#!/usr/bin/env python3
"""Stage 3 Task 3: assemble 청구항변반박전략서.md."""

from __future__ import annotations

import json
import re
from dataclasses import dataclass
from html import escape
from pathlib import Path
from typing import Any


BASE_DIR = Path(__file__).resolve().parent
CLIENT_GOAL_PATH = BASE_DIR / "client_goal.json"
DR_INFO_DIR = BASE_DIR / "Stage_3_Results" / "DR_Info_v2"
DEFENSE_DIR = BASE_DIR / "Stage_3_Results" / "항변반박_v2"
OUTPUT_PATH = BASE_DIR / "청구항변반박전략서.md"

CLAIM_FILE_RE = re.compile(r"^(C-(\d{3}))_info_D_R\.json$")
DEFENSE_TITLE_RE = re.compile(r"(?m)^##\s+청구 ID -\s*(C-\d{3})\s*-\s*항변/재반박\s*(\d+)\s*$")

FULL_WIDTH_TABLE_CSS = """<style>
.full-width-table-wrap {
  display:block !important;
  width:100% !important;
  min-width:100% !important;
  max-width:none !important;
  margin:0 !important;
  padding:0 !important;
  overflow:visible !important;
}
.full-width-table {
  width:100% !important;
  min-width:100% !important;
  max-width:none !important;
  table-layout:fixed !important;
  border-collapse:collapse !important;
  border-spacing:0 !important;
  margin:0 0 16px 0 !important;
  box-sizing:border-box !important;
}
.full-width-table th,
.full-width-table td {
  box-sizing:border-box !important;
  white-space:normal !important;
  word-break:break-word !important;
  overflow-wrap:anywhere !important;
}
</style>"""

TABLE_WRAP_STYLE = (
    "display:block; width:100%; min-width:100%; max-width:none; "
    "margin:0; padding:0; overflow:visible;"
)
TABLE_STYLE_BASE = (
    "width:100%; min-width:100%; max-width:none; table-layout:fixed; "
    "border-collapse:collapse; border-spacing:0; margin:0 0 16px 0; box-sizing:border-box;"
)
HEADER_CELL_STYLE = (
    "border:1px solid #444; padding:6px 8px; text-align:left; vertical-align:top; "
    "background:#f2f2f2; line-height:1.45; white-space:normal; word-break:break-word; overflow-wrap:anywhere;"
)
BODY_CELL_STYLE = (
    "border:1px solid #444; padding:6px 8px; text-align:left; vertical-align:top; "
    "line-height:1.5; white-space:normal; word-break:break-word; overflow-wrap:anywhere;"
)


@dataclass(frozen=True)
class ClaimInfo:
    claim_id: str
    claim_order: int
    claim_title: str
    case_kind: str
    plaintiffs: list[str]
    defendants: list[str]
    claim_statement: str
    relief_summary: str
    cause_summary: str
    facts: list[dict[str, Any]]
    legal_elements: list[str]


def read_json(path: Path) -> Any:
    with path.open("r", encoding="utf-8") as handle:
        return json.load(handle)


def read_text(path: Path) -> str:
    return path.read_text(encoding="utf-8")


def require_text(payload: dict[str, Any], key: str, source: Path) -> str:
    value = payload.get(key)
    if not isinstance(value, str) or not value.strip():
        raise ValueError(f"{source.name}: '{key}' must be a non-empty string.")
    return value.strip()


def require_string_list(payload: dict[str, Any], key: str, source: Path) -> list[str]:
    value = payload.get(key)
    if not isinstance(value, list):
        raise ValueError(f"{source.name}: '{key}' must be a list.")

    normalized: list[str] = []
    for item in value:
        if not isinstance(item, str) or not item.strip():
            raise ValueError(f"{source.name}: every item in '{key}' must be a non-empty string.")
        normalized.append(item.strip())
    return normalized


def load_claims() -> list[ClaimInfo]:
    claims: list[ClaimInfo] = []

    for path in sorted(DR_INFO_DIR.iterdir()):
        if not path.is_file():
            continue
        match = CLAIM_FILE_RE.match(path.name)
        if not match:
            continue

        payload = read_json(path)
        claim_id = require_text(payload, "claim_id", path)
        if claim_id != match.group(1):
            raise ValueError(f"{path.name}: claim_id '{claim_id}' does not match filename.")

        facts = payload.get("facts")
        if not isinstance(facts, list):
            raise ValueError(f"{path.name}: 'facts' must be a list.")

        legal_elements = require_string_list(payload, "legal_elements", path)

        claims.append(
            ClaimInfo(
                claim_id=claim_id,
                claim_order=int(match.group(2)),
                claim_title=require_text(payload, "claim_title", path),
                case_kind=require_text(payload, "case_kind", path),
                plaintiffs=require_string_list(payload, "plaintiffs", path),
                defendants=require_string_list(payload, "defendants", path),
                claim_statement=require_text(payload, "claim_statement", path),
                relief_summary=require_text(payload, "relief_summary", path),
                cause_summary=require_text(payload, "cause_summary", path),
                facts=facts,
                legal_elements=legal_elements,
            )
        )

    if not claims:
        raise FileNotFoundError(f"No claim files found in {DR_INFO_DIR}")

    return sorted(claims, key=lambda claim: claim.claim_order)


def normalize_text(text: str) -> str:
    return text.replace("\r\n", "\n").replace("\r", "\n")


def table_cell(value: Any) -> str:
    text = "" if value is None else str(value)
    text = normalize_text(text).strip()
    return text.replace("|", "\\|")


def format_block_text(text: str) -> str:
    normalized = normalize_text(text).strip()
    return re.sub(r"\n+", "<br>", normalized)


def format_legal_elements(elements: list[str]) -> str:
    return "<br>".join(table_cell(element) for element in elements)


def format_fact_actions(facts: list[dict[str, Any]]) -> str:
    lines: list[str] = []
    for fact in facts:
        fact_id = str(fact.get("fact_id", "")).strip()
        action = str(fact.get("action", "")).strip()
        if not fact_id or not action:
            continue
        lines.append(f"[{table_cell(fact_id)}]: {table_cell(action)}")
    return "<br>".join(lines)


def format_fact_evidence(facts: list[dict[str, Any]]) -> str:
    lines: list[str] = []
    for fact in facts:
        fact_id = str(fact.get("fact_id", "")).strip()
        raw_refs = fact.get("evidence_refs") or []
        if isinstance(raw_refs, str):
            refs = [raw_refs.strip()] if raw_refs.strip() else []
        elif isinstance(raw_refs, list):
            refs = [str(item).strip() for item in raw_refs if str(item).strip()]
        else:
            refs = []

        if not fact_id:
            continue

        evidence_text = ", ".join(table_cell(ref) for ref in refs) if refs else "증거 없음"
        lines.append(f"[{table_cell(fact_id)}]: {evidence_text}")
    return "<br>".join(lines)


def html_text(value: Any) -> str:
    text = "" if value is None else str(value)
    normalized = re.sub(r"\n+", "<br>", normalize_text(text).strip())
    parts = normalized.split("<br>")
    return "<br>".join(escape(part) for part in parts)


def render_html_table(
    headers: list[str],
    rows: list[list[str]],
    col_widths: list[int],
    *,
    font_size: str = "12px",
) -> str:
    if len(headers) != len(col_widths):
        raise ValueError("Header count and column width count must match.")

    table_style = f"{TABLE_STYLE_BASE} font-size:{font_size};"

    lines = [f'<div class="full-width-table-wrap" style="{TABLE_WRAP_STYLE}">']
    lines.append(f'  <table class="full-width-table" width="100%" style="{table_style}">')
    lines.append("    <colgroup>")
    for width in col_widths:
        lines.append(f'      <col style="width:{width}%;" />')
    lines.extend(["    </colgroup>", "    <thead>", "      <tr>"])

    for header in headers:
        lines.append(f'        <th style="{HEADER_CELL_STYLE}">{html_text(header)}</th>')

    lines.extend(["      </tr>", "    </thead>", "    <tbody>"])

    for row in rows:
        if len(row) != len(headers):
            raise ValueError("Every row must have the same number of columns as the header.")
        lines.append("      <tr>")
        for cell in row:
            lines.append(f'        <td style="{BODY_CELL_STYLE}">{html_text(cell)}</td>')
        lines.append("      </tr>")

    lines.extend(["    </tbody>", "  </table>", "</div>"])
    return "\n".join(lines)


def parse_markdown_row(line: str) -> list[str]:
    stripped = line.strip()
    if stripped.startswith("|"):
        stripped = stripped[1:]
    if stripped.endswith("|"):
        stripped = stripped[:-1]
    return [cell.strip() for cell in stripped.split("|")]


def parse_markdown_table(markdown_table: str) -> tuple[list[str], list[list[str]]]:
    table_lines = [line.strip() for line in normalize_text(markdown_table).splitlines() if line.strip().startswith("|")]
    if len(table_lines) < 3:
        raise ValueError("Markdown table must include a header, separator, and at least one data row.")

    headers = parse_markdown_row(table_lines[0])
    body_rows = [parse_markdown_row(line) for line in table_lines[2:]]
    return headers, body_rows


def collect_unique_parties(claims: list[ClaimInfo], field_name: str) -> list[str]:
    seen: set[str] = set()
    ordered: list[str] = []

    for claim in claims:
        for party in getattr(claim, field_name):
            if party in seen:
                continue
            seen.add(party)
            ordered.append(party)

    return ordered


def render_single_column_table(header: str, rows: list[str]) -> list[str]:
    lines = [f"| {header} |", "| --- |"]
    for row in rows:
        lines.append(f"| {table_cell(row)} |")
    return lines


def render_master_table(claims: list[ClaimInfo]) -> str:
    rows: list[list[str]] = []
    for claim in claims:
        plaintiffs = ", ".join(claim.plaintiffs)
        defendants = ", ".join(claim.defendants)
        rows.append(
            [
                claim.claim_id,
                claim.claim_title,
                claim.case_kind,
                plaintiffs,
                defendants,
            ]
        )
    return render_html_table(
        ["청구 ID", "청구권", "사건 종류", "원고", "피고"],
        rows,
        [10, 18, 16, 28, 28],
        font_size="11px",
    )


def render_claim_detail(claim: ClaimInfo) -> str:
    rows = [
        ["청구취지", format_block_text(claim.relief_summary)],
        ["청구원인", format_block_text(claim.cause_summary)],
        ["요건요소", format_legal_elements(claim.legal_elements)],
        ["법률행위", format_fact_actions(claim.facts)],
        ["서증 목록", format_fact_evidence(claim.facts)],
    ]
    return "\n".join(
        [
            f"#### {claim.claim_id}: {claim.claim_title}",
            "",
            render_html_table(["항목", "내용"], rows, [18, 82]),
        ]
    )


def render_defense_section(claim: ClaimInfo) -> str:
    defense_path = DEFENSE_DIR / f"{claim.claim_id}_defense_rebuttal.md"
    if not defense_path.exists():
        raise FileNotFoundError(f"Defense/rebuttal file not found: {defense_path}")

    raw_text = normalize_text(read_text(defense_path)).strip()
    raw_text = re.sub(r"(?m)^#\s+항변/재반박 매트릭스\s*$", "", raw_text).strip()

    matches = list(DEFENSE_TITLE_RE.finditer(raw_text))
    if len(matches) != 2:
        raise ValueError(f"{defense_path.name}: expected exactly 2 defense/rebuttal tables, found {len(matches)}.")

    claim_ids_in_doc = {match.group(1) for match in matches}
    if claim_ids_in_doc != {claim.claim_id}:
        raise ValueError(f"{defense_path.name}: unexpected claim IDs in defense markdown: {sorted(claim_ids_in_doc)}")

    rendered_sections: list[str] = []
    for index, match in enumerate(matches):
        table_number = match.group(2)
        start = match.end()
        end = matches[index + 1].start() if index + 1 < len(matches) else len(raw_text)
        table_block = raw_text[start:end].strip()
        headers, rows = parse_markdown_table(table_block)
        rendered_sections.extend(
            [
                f"### {claim.claim_id}: {claim.claim_title}, 항변/재반박 {table_number}",
                "",
                render_html_table(headers, rows, [22, 78]),
                "",
            ]
        )

    return "\n".join(rendered_sections).strip()


def build_document() -> str:
    client_goal = read_json(CLIENT_GOAL_PATH)
    primary_goal = require_text(client_goal, "primary_goal", CLIENT_GOAL_PATH)
    summary_key_incidents = require_string_list(client_goal, "summary_key_incidents", CLIENT_GOAL_PATH)

    claims = load_claims()
    plaintiffs = collect_unique_parties(claims, "plaintiffs")
    defendants = collect_unique_parties(claims, "defendants")

    lines: list[str] = [
        "# 청구/항변/반박 전략 문서",
        "",
        "## 1. 사건 개요",
        "### 1.1 원고 목표",
        primary_goal,
        "",
        "### 1.2 핵심 사실",
    ]
    lines.extend(f"- {incident}" for incident in summary_key_incidents)

    lines.extend(
        [
            "",
            "## 2. 당사자",
            "### 2.1 원고",
        ]
    )
    lines.extend(render_single_column_table("원고", plaintiffs))

    lines.extend(
        [
            "",
            "### 2.2 피고",
        ]
    )
    lines.extend(render_single_column_table("피고", defendants))

    lines.extend(
        [
            "",
            FULL_WIDTH_TABLE_CSS,
            "",
            "## 3. 청구 마스터 표",
            "### 3.1 청구권 개요",
            "",
        ]
    )
    lines.append(render_master_table(claims))

    lines.extend(
        [
            "",
            "### 3.2 개별 청구권 상세 정보",
            "",
        ]
    )

    for index, claim in enumerate(claims):
        if index:
            lines.append("")
        lines.append(render_claim_detail(claim))

    lines.extend(
        [
            "",
            "## 4. 예상항변 및 반박",
            "신뢰도는 현재 확보된 사실, 증거, 법리의 연결 정도를 바탕으로 한 상대평가이다. High는 바로 소장 및 준비서면 전략에 반영하기 좋은 수준, Medium은 기본 구조는 타당하지만 추가 사실 또는 증거 보강이 필요한 수준, Low는 방향성은 있으나 사실관계 또는 증거 연결의 추가 보강이 더 필요한 수준을 뜻한다.",
            "",
        ]
    )

    for claim in claims:
        lines.append(render_defense_section(claim))
        lines.append("")

    return "\n".join(lines).rstrip() + "\n"


def main() -> None:
    document = build_document()
    OUTPUT_PATH.write_text(document, encoding="utf-8")

    summary = {
        "status": "ok",
        "output_file": OUTPUT_PATH.as_posix(),
        "claim_count": len(load_claims()),
    }
    print(json.dumps(summary, ensure_ascii=False))


if __name__ == "__main__":
    main()
